06 - SFT 指令微调
预训练产出的叫 base 模型。这一篇把它变成能对话的模型,产出 sft.py。
前置:05 篇跑完的预训练 checkpoint,01 篇训好的 tokenizer。这一篇不需要多卡,0.5B 做 SFT 单卡就够。
零、开始之前:base 模型为什么不能直接用
0.1 base 模型学会了什么
预训练的目标从头到尾只有一个:猜下一个 token。所以 base 模型学到的是这段文本后面接什么最自然。
它读了 10B token 的网页,语法、常识、写作风格都有了。但它不知道「有人在问你问题,你该回答」这回事,因为训练数据里从来没有这种结构。
0.2 拿它当聊天机器人会发生什么
给 base 模型输入「中国的首都是哪里?」,它可能这样接:
中国的首都是哪里?
A. 上海
B. 北京
C. 广州
D. 深圳
答案:B
解析:北京是中华人民共和国的首都……
它没有回答你,它在续写。因为在网页语料里,这样一句话后面最常见的就是选择题选项。
也可能接出这样的:
中国的首都是哪里?这是很多小学生都会问的问题。今天我们就来聊聊……
同样是续写,不是回答。
这就是 base 模型的状态:知识都在,但不知道自己该扮演一个回答问题的角色。
0.3 SFT 要解决的是格式问题,不是知识问题
这一点特别容易误解,说清楚:
SFT 教的是行为模式,不是知识。 模型的知识在预训练阶段就基本定型了,SFT 用的那几万条数据,相对 10B token 来说是九牛一毛,不可能塞进什么新知识。
SFT 让模型学会的是:看到 <|im_start|>user 开头的一段,后面跟 <|im_start|>assistant 时,我应该输出一个回答,然后停下来。
这也解释了一个常见现象:SFT 数据里出现的知识错误,模型会照单全收;但想靠 SFT 给模型灌输一个新领域的知识,基本没用。 前者是学格式时顺带记住了内容,后者是想用几万条覆盖百亿 token 建立的分布。
有篇叫 LIMA 的论文把这个观点推到极致,它只用 1000 条精心挑选的数据做 SFT,效果就不错。结论是:SFT 数据的质量远比数量重要,因为你本来就只是在教格式。
0.4 跟预训练的三处不同
| 预训练 | SFT | |
|---|---|---|
| 数据组织 | 所有文档首尾拼成一条 token 流 | 按样本组织,要 padding |
| loss 算在哪 | 每个 token 都算 | 只在 assistant 的回复上算 |
| 学习率 | 3e-4 | 2e-5,小一个数量级 |
| 跑多久 | 一遍 10B token | 同样的数据跑 2 到 3 轮 |
这一篇剩下的内容就是把这四行讲清楚。
一、对话模板
1.1 为什么需要模板
模型的输入是 一串 token,本身没有「谁在说话」的概念。要让它区分用户和助手,必须在文本里用特殊标记把角色写出来。
这套标记方式就是对话模板。用哪套都行,但训练和推理必须用同一套,否则模型认不出来。
1.2 ChatML
我们用 ChatML,现在最通行的一种,Qwen、多数开源模型都用它:
<|im_start|>system
你是一个乐于助人的助手。<|im_end|>
<|im_start|>user
中国的首都是哪里?<|im_end|>
<|im_start|>assistant
北京。<|im_end|>
规则很简单:每一轮是 <|im_start|> + 角色名 + 换行 + 内容 + <|im_end|>。
推理的时候,把用户输入按这个格式拼好,最后补上 <|im_start|>assistant\n,模型就会接着往下生成回复,生成到 <|im_end|> 就停。
1.3 special token 现在派上用场了
<|im_start|> 和 <|im_end|> 就是 01 篇 4.8 节训 tokenizer 时留的那两个。
当时说「预训练根本用不上,但现在不加,将来加就得重训」,现在到了要用的时候。如果那时候没留,此刻加进去词表会从 32000 变成 32002,embedding 层形状对不上,预训练的权重全废。
代码里加了个断言,就是防这个:
assert self.im_start is not None and self.im_end is not None, (
"tokenizer 里没有 <|im_start|> / <|im_end|>,"
"回 01 篇 4.8 节,这两个必须在训 tokenizer 时就加进 special_tokens"
)
1.4 为什么要把这两个当成单个 token
如果不把 <|im_start|> 注册成 special token,BPE 会把它切成 <、|、im、_、start、|、> 一堆碎片。
坏处有两个。一是浪费,每轮对话多花十几个 token。二是不可靠,模型要学会「这七个碎片连在一起才表示角色开始」,比学一个独立符号难得多,而且用户在正常文本里打出 <|im_start|> 就能伪造角色,这是提示注入的一个口子。
注册成单个 token 之后,它在词表里有独立编号,跟任何正常文本都撞不上。
二、loss mask:这一篇最核心的改动
2.1 不该在用户的话上算 loss
一条 SFT 样本里有三部分:模板标记、用户的问题、助手的回答。
预训练是每个 token 都算 loss。但 SFT 不能这样,只有助手的回答该算。
2.2 为什么
在用户的问题上算 loss,等于在教模型「怎么提问」。
后果是模型学会了模仿用户的说话方式,推理时容易自问自答:你问一句,它回答完接着又冒出一个新问题然后自己回答。这是 SFT 最典型的翻车形式之一。
而且这部分 loss 是纯噪声。用户会问什么是不可预测的,让模型去拟合它没有任何意义,只会稀释真正有用的梯度信号。
2.3 怎么做
PyTorch 的 F.cross_entropy 有个 ignore_index 参数,标成这个值的位置直接跳过。约定俗成用 -100:
IGNORE = -100
构造 labels 时,非助手回复的位置全填 -100:
im_end 必须算进去,它是模型学会「何时停止」的唯一信号。这类错误不会报错、loss 照样下降,只能靠 2.6 节把 mask 打印出来肉眼核对。输入 <|im_start|>user \n 中国的首都是哪里? <|im_end|> \n <|im_start|>assistant \n 北京。 <|im_end|>
labels -100 -100 -100 -100... -100 -100 -100 -100 -100 北京。 <|im_end|>
2.4 代码
def encode(self, messages: list) -> tuple:
ids, labels = [], []
for m in messages:
role = self.tok.encode(m["role"]).ids
content = self.tok.encode(m["content"]).ids
# 头部 <|im_start|>{role}\n 不算 loss
head = [self.im_start] + role + self.nl
ids += head
labels += [IGNORE] * len(head)
body = content + [self.im_end] + self.nl
ids += body
if m["role"] == "assistant":
labels += body # 只有这一支真的填 labels
else:
labels += [IGNORE] * len(body)
return ids, labels
注意 <|im_end|> 是算 loss 的。 这一点很关键:模型必须学会「回答完了要输出结束符」,否则推理时它会一直说下去停不下来。如果把 <|im_end|> 也 mask 掉,训出来的模型会滔滔不绝,这是另一种典型翻车。
2.5 算 loss 时还要错位一格
跟预训练一样,用第 t 个位置的输出预测第 t+1 个 token,所以要错位:
def sft_loss(logits, labels):
shift_logits = logits[:, :-1, :].contiguous()
shift_labels = labels[:, 1:].contiguous()
return F.cross_entropy(shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1),
ignore_index=IGNORE)
预训练时这个错位是在 get_batch 里做的(y 比 x 右移一位),SFT 因为要配合 mask,在算 loss 时做更清楚。
2.6 必须把 mask 打印出来看一眼
这是这一篇最该照做的一条建议。loss mask 写错了不会报错,训练照跑,loss 照降,只是降的是错的东西。
所以写完先可视化验证:
def inspect_masking(dataset, tok, n=2):
for i in range(min(n, len(dataset))):
ids, labels = dataset[i]
kept = [t for t, l in zip(ids, labels) if l != IGNORE]
print("完整输入:", tok.decode(ids)[:300])
print("算 loss 的部分:", tok.decode(kept)[:300])
print(f"占比 {len(kept)}/{len(ids)} = {len(kept) / len(ids):.1%}")
python sft.py --ckpt out/ckpt.pt --data data/sft.jsonl --inspect_only
「算 loss 的部分」打印出 来应该只有助手的回答加上结束符,一个字的用户输入都不该有。
顺便看那个占比。正常在 30% 到 60% 之间。明显偏高说明 mask 没生效,明显偏低说明可能把回复也 mask 掉了。
三、数据组织
3.1 不能再首尾拼接
01 篇 5.3 节说过,预训练把所有文档拼成一条流,随机取起点,不管文档边界。SFT 不行。
原因是一条对话必须完整。截断了的样本,模型看到半句问题就要生成回答,学到的是错的模式。而且 loss mask 是按样本结构算的,拼接之后结构就乱了。
所以 SFT 回到常规做法:一条样本一条,不够长的 padding 补齐。
3.2 右侧 padding 配因果注意力,不需要 padding mask
这里有个值得讲清楚的细节。padding 补进去的位置是垃圾数据,直觉上应该用 attention mask 屏蔽掉。但训练时用右侧 padding 的话,其实不需要。
理由两条:
真实 token 在前,padding 在后。02 篇的因果掩码保证每个位置只能看到自己和前面,所以真实 token 永远不会注意到后面的 padding。
padding 位置自己的输出确实是垃圾,但它们在 labels 里是 IGNORE,不进 loss。
所以右侧 padding 加因果注意力,padding 天然就被隔离了。推理时则相反:
什么时候需要 attention mask:推理时批量生成要用左侧 padding(因为生成是从序列末尾往后接,末尾必须对齐),那时候 padding 在前面,真实 token 会注意到它们,必须显式屏蔽。
代码里 collate 仍然返回了 attn,但训练循环没用它,就是为了让这个区别显式可见。
def collate(batch, pad_id: int):
maxlen = max(len(ids) for ids, _ in batch)
input_ids, labels, attn = [], [], []
for ids, lab in batch:
pad = maxlen - len(ids)
input_ids.append(ids + [pad_id] * pad)
labels.append(lab + [IGNORE] * pad)
attn.append([1] * len(ids) + [0] * pad)
return torch.tensor(input_ids), torch.tensor(labels), torch.tensor(attn)
padding 到本 batch 内的最长,不是 padding 到 max_len。差别很大:按 batch 内最长补,平均只多算百分之几十;按 2048 补,短样本会浪费十几倍算力。
3.3 超长样本直接丢,不要截断
if len(ids) > max_len:
skipped += 1
continue
截断会把助手的回答砍掉一半,模型学到「回答可以不说完」,比丢掉这条样本糟糕得多。
统计跳过了多少条,如果比例很高(超过百分之几),说明 max_len 设小了。
四、指令数据从哪来
4.1 几个常用的开源数据集
以下行数和体积是 2026-08-19 用 HuggingFace datasets-server 实测的:
| 数据集 | 行数 | 体积 | 许可 | 说明 |
|---|---|---|---|---|
HuggingFaceH4/ultrachat_200k | 515,311 | 1,624.0 MB | MIT | 多轮对话,质量高,英文 |
shibing624/sharegpt_gpt4 | 103,415 | 602.3 MB | CC-BY-4.0 | GPT-4 生成的中文多轮 |
tatsu-lab/alpaca | 52,002 | 24.2 MB | CC-BY-NC-4.0 | 经典单轮,注意是非商用 |
llamafactory/alpaca_gpt4_zh | 42,677 | 27.9 MB | Apache-2.0 | Alpaca 的 GPT-4 中文版 |
GAIR/lima | 1,000 | — | other | LIMA 论文那 1000 条,质量标杆 |
选的时候注意许可证。tatsu-lab/alpaca 是 CC-BY-NC,只能非商用,很多人没留意就用了。
4.2 要多少条
按 0.3 节说的,SFT 是教格式,几万条足够。
建议从两三万条混合数据起步,中英比例对齐 tokenizer 的语料配比。不足再加,不要一次堆几十万条。
理由是 SFT 阶段调试的循环要快。0.5B 模型两万条数据跑 3 轮,单卡几十分钟就完事,能快速看到效果、快速调整。堆到几十万条,一轮就要好几小时,反馈太慢。
4.3 数据格式
统一成 jsonl,每行一条:
{"messages": [{"role": "user", "content": "中国的首都是哪里?"}, {"role": "assistant", "content": "北京。"}]}
各个开源数据集的原始格式五花八门(Alpaca 是 instruction/input/output,ShareGPT 是 conversations 数组),先写个转换脚本统一成上面这种,后面的代码才好写。
五、超参怎么变
| 超参 | 预训练 | SFT | 为什么 |
|---|---|---|---|
| 学习率 | 3e-4 | 2e-5 | 小一个数量级 |
| epoch | 相当于不到 1 轮 | 2 到 3 轮 | 数据少,要多看几遍 |
| warmup | 2% | 3% | 数据少,比例稍高 |
| weight decay | 0.1 | 0 | 不需要正则 |
| batch size | 1M token | 十几条样本 | 数据量差三个数量级 |
学习率是这里最要紧的。 预训练是从随机初始化开始学,要大步走。SFT 是在一个已经学好的模型上做微调,目标只是调整输出格式,大学习率会把预训练学到的能力冲掉,这个现象叫灾难性遗忘。表现是模型学会了对话格式,但一问知识就胡说八道。
2e-5 是通行值。不确定的话宁可再小一点,1e-5 也行。
epoch 不要多。 数据只有几万条,跑太多轮会过拟合,模型开始逐字背诵训练数据。一般 2 到 3 轮,边跑边看验证集,效果不涨就停。
六、LoRA
6.1 全参微调的问题
上面讲的都是全参微调,所有 502M 参数都参与更新。对 0.5B 来说这完全可行,02 篇算过静态显存 8 GB。
但如果基座换成 7B、13B,全参微调的显存就吃不消了,而且每个下游任务都要存一份完整的模型权重。
6.2 LoRA 的想法
LoRA 的观察是:微调对权重的改动量 ,其实是低秩的。也就是说 虽然是个大矩阵,但它包含的有效信息很少,可以用两个瘦长矩阵的乘积近似。
于是做法变成:冻结原权重 ,在旁边加一条支路:
其中 远小于 和 ,叫秩。训练时只更新 和 。
用随机初始化, 初始化成全 0,这样训练开始时 ,模型输出和原来完全一样,是个平滑的起点。
推理时可以把 加回 ,合并成一个矩阵,推理速度和原模型完全一样,没有额外开销。这是 LoRA 相比其他微调方法的一个重要优点。
6.3 参数量算一算
原矩阵是 ,LoRA 是